diff --git a/backend/README.md b/backend/README.md index 35275b5..a0e495c 100644 --- a/backend/README.md +++ b/backend/README.md @@ -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 文档: 团队接口清单见 `../docs/后端接口契约-开发版.md`,机器可读契约以运行时的 `/openapi.json` 为准。 + +AI Core 与 Agent Core 的模块边界、Mock Provider 和 Tool Calling 调试方式见 `../docs/AI-Core与Agent-Core开发说明.md`。 diff --git a/backend/app/agent/__init__.py b/backend/app/agent/__init__.py new file mode 100644 index 0000000..d860342 --- /dev/null +++ b/backend/app/agent/__init__.py @@ -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", +] diff --git a/backend/app/agent/builtin_tools.py b/backend/app/agent/builtin_tools.py new file mode 100644 index 0000000..2ffafbf --- /dev/null +++ b/backend/app/agent/builtin_tools.py @@ -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, + ) diff --git a/backend/app/agent/permissions.py b/backend/app/agent/permissions.py new file mode 100644 index 0000000..e28cb9e --- /dev/null +++ b/backend/app/agent/permissions.py @@ -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) diff --git a/backend/app/agent/runtime.py b/backend/app/agent/runtime.py new file mode 100644 index 0000000..08bf61c --- /dev/null +++ b/backend/app/agent/runtime.py @@ -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 diff --git a/backend/app/agent/tools.py b/backend/app/agent/tools.py new file mode 100644 index 0000000..103b54f --- /dev/null +++ b/backend/app/agent/tools.py @@ -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), + ) diff --git a/backend/app/container.py b/backend/app/container.py new file mode 100644 index 0000000..3377e2d --- /dev/null +++ b/backend/app/container.py @@ -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() diff --git a/backend/app/contracts.py b/backend/app/contracts.py index 88eb18b..c86eb2d 100644 --- a/backend/app/contracts.py +++ b/backend/app/contracts.py @@ -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" diff --git a/backend/app/main.py b/backend/app/main.py index 0e97f93..510e019 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -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( diff --git a/backend/app/providers/__init__.py b/backend/app/providers/__init__.py new file mode 100644 index 0000000..0555af2 --- /dev/null +++ b/backend/app/providers/__init__.py @@ -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", +] diff --git a/backend/app/providers/base.py b/backend/app/providers/base.py new file mode 100644 index 0000000..039d3f3 --- /dev/null +++ b/backend/app/providers/base.py @@ -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]: ... diff --git a/backend/app/providers/mock.py b/backend/app/providers/mock.py new file mode 100644 index 0000000..381641d --- /dev/null +++ b/backend/app/providers/mock.py @@ -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." diff --git a/backend/app/providers/registry.py b/backend/app/providers/registry.py new file mode 100644 index 0000000..c19e7a0 --- /dev/null +++ b/backend/app/providers/registry.py @@ -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, + ) diff --git a/backend/app/routes.py b/backend/app/routes.py index 2d2b7c3..b847d12 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -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 diff --git a/backend/tests/test_agent_core.py b/backend/tests/test_agent_core.py new file mode 100644 index 0000000..0aa843b --- /dev/null +++ b/backend/tests/test_agent_core.py @@ -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()) diff --git a/backend/tests/test_api.py b/backend/tests/test_api.py index 8c8dc64..d1a17fc 100644 --- a/backend/tests/test_api.py +++ b/backend/tests/test_api.py @@ -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" diff --git a/docs/AI-Core与Agent-Core开发说明.md b/docs/AI-Core与Agent-Core开发说明.md new file mode 100644 index 0000000..1c88dfc --- /dev/null +++ b/docs/AI-Core与Agent-Core开发说明.md @@ -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。 diff --git a/docs/后端接口契约-开发版.md b/docs/后端接口契约-开发版.md index 78cf55d..01a59df 100644 --- a/docs/后端接口契约-开发版.md +++ b/docs/后端接口契约-开发版.md @@ -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 逻辑。