添加provider工厂和Ollama支持
- 实现ProviderFactory用于构建不同类型的provider适配器 - 添加EnvironmentCredentialResolver用于解析环境变量中的凭证 - 实现OllamaProvider支持本地模型调用 - 实现OpenAICompatibleProvider支持OpenAI兼容接口 - 在AgentRuntime中添加对ProviderError的处理 - 更新Message结构体添加tool_calls字段 - 实现provider配置的增删改查API端点 - 添加provider注册表的replace方法 - 添加HTTP基础类和工具参数解码功能 - 更新依赖添加httpx库 - 添加相关单元测试验证provider适配器功能 ```
This commit is contained in:
@@ -2,6 +2,8 @@ import asyncio
|
||||
|
||||
from app.main import health, service_status
|
||||
from app.routes import get_index_status, list_notes, list_plugins, list_providers, list_skills
|
||||
from app.routes import create_provider, delete_provider, get_provider, update_provider
|
||||
from app.contracts import ProviderCreateRequest, ProviderType, ProviderUpdateRequest
|
||||
|
||||
|
||||
def test_health() -> None:
|
||||
@@ -53,3 +55,25 @@ def test_openapi_contains_documented_frontend_interfaces() -> None:
|
||||
}
|
||||
|
||||
assert expected_paths <= paths.keys()
|
||||
|
||||
|
||||
def test_provider_configuration_lifecycle() -> None:
|
||||
created = asyncio.run(
|
||||
create_provider(
|
||||
ProviderCreateRequest(
|
||||
provider_type=ProviderType.ollama,
|
||||
name="Local Ollama",
|
||||
base_url="http://127.0.0.1:11434",
|
||||
default_model="qwen3:latest",
|
||||
)
|
||||
)
|
||||
)
|
||||
fetched = asyncio.run(get_provider(created.provider_id))
|
||||
disabled = asyncio.run(
|
||||
update_provider(created.provider_id, ProviderUpdateRequest(enabled=False))
|
||||
)
|
||||
deleted = asyncio.run(delete_provider(created.provider_id))
|
||||
|
||||
assert fetched.provider_type == ProviderType.ollama
|
||||
assert disabled.enabled is False
|
||||
assert deleted.resource_id == created.provider_id
|
||||
|
||||
@@ -0,0 +1,154 @@
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
import httpx
|
||||
|
||||
from app.contracts import Message, MessageRole, ModelRequest, ToolCall, ToolDefinition
|
||||
from app.providers.ollama import OllamaProvider
|
||||
from app.providers.openai_compatible import OpenAICompatibleProvider
|
||||
|
||||
|
||||
class StaticCredentials:
|
||||
def resolve(self, credential_id: str | None) -> str | None:
|
||||
return "secret-test-key" if credential_id else None
|
||||
|
||||
|
||||
def run(coroutine):
|
||||
return asyncio.run(coroutine)
|
||||
|
||||
|
||||
def test_openai_compatible_maps_tool_call_and_credentials() -> None:
|
||||
captured: dict = {}
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
assert request.headers["Authorization"] == "Bearer secret-test-key"
|
||||
captured.update(json.loads(request.content))
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "math.add",
|
||||
"arguments": '{"left":1,"right":2}',
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 8, "completion_tokens": 4},
|
||||
},
|
||||
)
|
||||
|
||||
provider = OpenAICompatibleProvider(
|
||||
base_url="https://provider.test/v1",
|
||||
credential_id="openai-test",
|
||||
credentials=StaticCredentials(),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
turn = run(
|
||||
provider.complete(
|
||||
ModelRequest(
|
||||
provider_id="test",
|
||||
model="test-model",
|
||||
messages=[Message(role=MessageRole.user, content="add")],
|
||||
tools=[
|
||||
ToolDefinition(
|
||||
name="math.add",
|
||||
description="Add numbers",
|
||||
parameters={"type": "object"},
|
||||
)
|
||||
],
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
assert captured["tools"][0]["function"]["name"] == "math.add"
|
||||
assert turn.tool_calls[0].name == "math.add"
|
||||
assert turn.tool_calls[0].arguments == {"left": 1, "right": 2}
|
||||
assert turn.input_tokens == 8
|
||||
|
||||
|
||||
def test_openai_compatible_preserves_tool_call_context() -> None:
|
||||
captured: dict = {}
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured.update(json.loads(request.content))
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"choices": [{"message": {"role": "assistant", "content": "done"}}],
|
||||
"usage": {},
|
||||
},
|
||||
)
|
||||
|
||||
provider = OpenAICompatibleProvider(
|
||||
base_url="https://provider.test/v1",
|
||||
credential_id=None,
|
||||
credentials=StaticCredentials(),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
call = ToolCall(
|
||||
tool_call_id="call_1", name="math.add", arguments={"left": 1, "right": 2}
|
||||
)
|
||||
turn = run(
|
||||
provider.complete(
|
||||
ModelRequest(
|
||||
provider_id="test",
|
||||
model="test-model",
|
||||
messages=[
|
||||
Message(role=MessageRole.user, content="add"),
|
||||
Message(role=MessageRole.assistant, content="", tool_calls=[call]),
|
||||
Message(
|
||||
role=MessageRole.tool,
|
||||
content='{"value":3}',
|
||||
name="math.add",
|
||||
tool_call_id="call_1",
|
||||
),
|
||||
],
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
assert captured["messages"][1]["tool_calls"][0]["id"] == "call_1"
|
||||
assert captured["messages"][2]["tool_call_id"] == "call_1"
|
||||
assert turn.text == "done"
|
||||
|
||||
|
||||
def test_ollama_maps_models_and_completion() -> None:
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
if request.url.path == "/api/tags":
|
||||
return httpx.Response(200, json={"models": [{"name": "qwen3:latest"}]})
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"message": {"role": "assistant", "content": "local answer"},
|
||||
"prompt_eval_count": 5,
|
||||
"eval_count": 2,
|
||||
},
|
||||
)
|
||||
|
||||
provider = OllamaProvider(transport=httpx.MockTransport(handler))
|
||||
models = run(provider.list_models())
|
||||
turn = run(
|
||||
provider.complete(
|
||||
ModelRequest(
|
||||
provider_id="ollama",
|
||||
model="qwen3:latest",
|
||||
messages=[Message(role=MessageRole.user, content="hello")],
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
assert models[0].model == "qwen3:latest"
|
||||
assert turn.text == "local answer"
|
||||
assert turn.input_tokens == 5
|
||||
assert turn.output_tokens == 2
|
||||
Reference in New Issue
Block a user