Files

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