实现 AI Core 与 Agent Core 基础功能
- 更新 README 描述从后端壳子到 AI Core/Agent Core - 添加 ToolCall 和 ToolResult 数据结构定义 - 扩展 AgentRun 模型增加输出、错误码、工具调用结果等字段 - 添加 mock 提供商类型支持 - 实现聊天、代理运行、工具调用和提供商管理的核心路由逻辑 - 集成容器化依赖注入和错误处理机制 - 更新 API 接口契约和文档说明
This commit is contained in:
@@ -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,
|
||||
)
|
||||
Reference in New Issue
Block a user