feat: add AI research chatbot with confirmed business tools
This commit is contained in:
@@ -0,0 +1,117 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user