118 lines
5.2 KiB
Python
118 lines
5.2 KiB
Python
"""Provider creation and synthetic capability checks, with no business access."""
|
|
|
|
from contextlib import asynccontextmanager
|
|
|
|
import httpx
|
|
from openai import AsyncOpenAI
|
|
from pydantic_ai import Agent, AgentRunResultEvent
|
|
from pydantic_ai.messages import PartDeltaEvent, PartStartEvent, TextPart, TextPartDelta
|
|
from pydantic_ai.models.openai import OpenAIChatModel, OpenAIResponsesModel
|
|
from pydantic_ai.profiles.openai import OpenAIModelProfile
|
|
from pydantic_ai.providers.openai import OpenAIProvider
|
|
from pydantic_ai.usage import UsageLimits
|
|
|
|
from ..security import cipher
|
|
|
|
|
|
@asynccontextmanager
|
|
async def model_connection(config, settings, transport=None):
|
|
"""Connect only to the administrator's saved endpoint; never follow key-bearing redirects."""
|
|
key = cipher(settings).decrypt(config.api_key_encrypted.encode()).decode()
|
|
async with httpx.AsyncClient(
|
|
transport=transport, follow_redirects=False, trust_env=False, timeout=settings.ai_timeout
|
|
) as http:
|
|
client = AsyncOpenAI(base_url=config.base_url, api_key=key, http_client=http, max_retries=0)
|
|
provider = OpenAIProvider(openai_client=client)
|
|
cls = OpenAIChatModel if config.protocol == "chat_completions" else OpenAIResponsesModel
|
|
# Schema validation belongs to our backend even when a compatible gateway lacks strict mode.
|
|
model = cls(
|
|
config.model,
|
|
provider=provider,
|
|
profile=OpenAIModelProfile(openai_supports_strict_tool_definition=False),
|
|
)
|
|
yield model
|
|
|
|
|
|
def public_error(exc):
|
|
"""Never return provider response bodies, URLs, keys, or SDK exception strings."""
|
|
code = getattr(exc, "status_code", None)
|
|
if code in (401, 403):
|
|
return "模型服务拒绝访问,请检查 API Key 和模型权限"
|
|
if code == 404:
|
|
return "模型或接口不存在,请检查 Base URL、模型标识及接口协议"
|
|
if code == 429:
|
|
return "模型服务限流或额度不足,请稍后重试"
|
|
if isinstance(exc, (TimeoutError, httpx.TimeoutException)):
|
|
return "模型服务响应超时,请重试或检查服务状态"
|
|
return "模型服务调用失败,请检查连接和接口兼容性"
|
|
|
|
|
|
def ensure_complete(model, result):
|
|
"""A closed socket without a provider terminal frame is an interrupted stream."""
|
|
from pydantic_ai.messages import ModelResponse
|
|
|
|
if isinstance(model, (OpenAIChatModel, OpenAIResponsesModel)):
|
|
responses = [m for m in result.new_messages() if isinstance(m, ModelResponse)]
|
|
if not responses or any(m.finish_reason is None for m in responses):
|
|
raise ValueError("Provider stream did not include a completion marker")
|
|
|
|
|
|
async def test_capabilities(model):
|
|
"""Require actual stream text plus a test tool call followed by its exact output."""
|
|
from uuid import uuid4
|
|
|
|
result = {}
|
|
try:
|
|
streamed = False
|
|
text = ""
|
|
async with Agent(model, tool_retries=0, output_retries=0).run_stream_events(
|
|
"Reply with READY.", model_settings={"max_tokens": 256}, usage_limits=UsageLimits(request_limit=1)
|
|
) as stream:
|
|
async for event in stream:
|
|
if isinstance(event, AgentRunResultEvent):
|
|
ensure_complete(model, event.result)
|
|
if isinstance(event, PartStartEvent) and isinstance(event.part, TextPart):
|
|
streamed = True
|
|
text += event.part.content
|
|
elif isinstance(event, PartDeltaEvent) and isinstance(event.delta, TextPartDelta):
|
|
streamed = True
|
|
text += event.delta.content_delta
|
|
result["answer"] = {
|
|
"ok": bool(text.strip()),
|
|
"message": "收到回答" if text.strip() else "没有收到文本回答",
|
|
}
|
|
result["stream"] = {
|
|
"ok": streamed and bool(text.strip()),
|
|
"message": "收到流式文本" if streamed else "没有收到流式文本",
|
|
}
|
|
except Exception as exc:
|
|
result["answer"] = result["stream"] = {"ok": False, "message": public_error(exc)}
|
|
secret = uuid4().hex
|
|
called = False
|
|
|
|
async def capability_probe() -> str:
|
|
"""Read a random test marker. Call this tool and repeat its returned marker exactly."""
|
|
nonlocal called
|
|
called = True
|
|
return secret
|
|
|
|
try:
|
|
final = None
|
|
async with Agent(model, tools=[capability_probe], tool_retries=0, output_retries=0).run_stream_events(
|
|
"Call capability_probe, then reply with the exact marker it returned. Do not guess.",
|
|
model_settings={"max_tokens": 256},
|
|
usage_limits=UsageLimits(request_limit=2, tool_calls_limit=1),
|
|
) as stream:
|
|
async for event in stream:
|
|
if isinstance(event, AgentRunResultEvent):
|
|
ensure_complete(model, event.result)
|
|
final = event.result.output
|
|
ok = called and isinstance(final, str) and secret in final
|
|
result["tools"] = {
|
|
"ok": ok,
|
|
"message": "工具调用与结果回传成功" if ok else "工具调用或结果回传未通过",
|
|
}
|
|
except Exception as exc:
|
|
result["tools"] = {"ok": False, "message": public_error(exc)}
|
|
return result
|