795 lines
27 KiB
Python
795 lines
27 KiB
Python
"""
|
||
LLM 提供方测试服务
|
||
"""
|
||
import asyncio
|
||
import json
|
||
import socket
|
||
import time
|
||
import urllib.error
|
||
import urllib.parse
|
||
import urllib.request
|
||
import ssl
|
||
from typing import Any, Dict, List, Optional
|
||
|
||
|
||
PROVIDER_CATALOG: List[Dict[str, str]] = [
|
||
{
|
||
"value": "openai",
|
||
"label": "OpenAI",
|
||
"default_endpoint_url": "https://api.openai.com/v1",
|
||
"protocol": "openai_compatible",
|
||
},
|
||
{
|
||
"value": "deepseek",
|
||
"label": "DeepSeek",
|
||
"default_endpoint_url": "https://api.deepseek.com/v1",
|
||
"protocol": "openai_compatible",
|
||
},
|
||
{
|
||
"value": "anthropic",
|
||
"label": "Anthropic",
|
||
"default_endpoint_url": "https://api.anthropic.com",
|
||
"protocol": "anthropic",
|
||
},
|
||
{
|
||
"value": "gemini",
|
||
"label": "Google Gemini",
|
||
"default_endpoint_url": "https://generativelanguage.googleapis.com/v1beta",
|
||
"protocol": "gemini",
|
||
},
|
||
{
|
||
"value": "dashscope",
|
||
"label": "阿里百炼",
|
||
"default_endpoint_url": "https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||
"protocol": "openai_compatible",
|
||
},
|
||
{
|
||
"value": "zhipu",
|
||
"label": "智谱 AI",
|
||
"default_endpoint_url": "https://open.bigmodel.cn/api/paas/v4",
|
||
"protocol": "openai_compatible",
|
||
},
|
||
{
|
||
"value": "moonshot",
|
||
"label": "Moonshot AI",
|
||
"default_endpoint_url": "https://api.moonshot.cn/v1",
|
||
"protocol": "openai_compatible",
|
||
},
|
||
{
|
||
"value": "groq",
|
||
"label": "Groq",
|
||
"default_endpoint_url": "https://api.groq.com/openai/v1",
|
||
"protocol": "openai_compatible",
|
||
},
|
||
{
|
||
"value": "openrouter",
|
||
"label": "OpenRouter",
|
||
"default_endpoint_url": "https://openrouter.ai/api/v1",
|
||
"protocol": "openai_compatible",
|
||
},
|
||
{
|
||
"value": "siliconflow",
|
||
"label": "SiliconFlow",
|
||
"default_endpoint_url": "https://api.siliconflow.cn/v1",
|
||
"protocol": "openai_compatible",
|
||
},
|
||
{
|
||
"value": "ollama",
|
||
"label": "Ollama",
|
||
"default_endpoint_url": "http://localhost:11434/v1",
|
||
"protocol": "openai_compatible",
|
||
},
|
||
{
|
||
"value": "ark",
|
||
"label": "火山方舟",
|
||
"default_endpoint_url": "https://ark.cn-beijing.volces.com/api/v3",
|
||
"protocol": "openai_compatible",
|
||
},
|
||
{
|
||
"value": "custom",
|
||
"label": "自定义兼容接口",
|
||
"default_endpoint_url": "",
|
||
"protocol": "openai_compatible",
|
||
},
|
||
]
|
||
|
||
PROVIDER_MAP = {item["value"]: item for item in PROVIDER_CATALOG}
|
||
|
||
|
||
class LLMProviderService:
|
||
"""LLM 提供方测试服务"""
|
||
|
||
@staticmethod
|
||
def get_provider_catalog() -> List[Dict[str, str]]:
|
||
return PROVIDER_CATALOG
|
||
|
||
@staticmethod
|
||
def get_provider_label(provider: Optional[str]) -> str:
|
||
if not provider:
|
||
return "自定义模型"
|
||
return PROVIDER_MAP.get(provider, {}).get("label", provider)
|
||
|
||
@staticmethod
|
||
def get_default_endpoint_url(provider: Optional[str]) -> str:
|
||
if not provider:
|
||
return ""
|
||
return PROVIDER_MAP.get(provider, {}).get("default_endpoint_url", "")
|
||
|
||
@classmethod
|
||
def build_model_name(cls, provider: Optional[str], llm_model_name: str) -> str:
|
||
model_name = (llm_model_name or "").strip()
|
||
label = cls.get_provider_label(provider)
|
||
if not model_name:
|
||
return label
|
||
return f"{label} {model_name}"
|
||
|
||
@staticmethod
|
||
def build_model_code(provider: Optional[str], llm_model_name: str) -> str:
|
||
provider_part = (provider or "custom").strip().lower()
|
||
model_part = (llm_model_name or "").strip().lower()
|
||
sanitized = []
|
||
previous_is_separator = False
|
||
for char in model_part:
|
||
if char.isalnum():
|
||
sanitized.append(char)
|
||
previous_is_separator = False
|
||
else:
|
||
if not previous_is_separator:
|
||
sanitized.append("_")
|
||
previous_is_separator = True
|
||
|
||
model_slug = "".join(sanitized).strip("_") or "model"
|
||
return f"llm_{provider_part}_{model_slug}"
|
||
|
||
@classmethod
|
||
async def test_model_connection(cls, payload: Dict[str, Any]) -> Dict[str, Any]:
|
||
provider = payload.get("provider")
|
||
provider_meta = PROVIDER_MAP.get(provider, PROVIDER_MAP["custom"])
|
||
endpoint_url = (payload.get("endpoint_url") or provider_meta.get("default_endpoint_url") or "").strip()
|
||
llm_model_name = (payload.get("llm_model_name") or "").strip()
|
||
api_key = (payload.get("api_key") or "").strip()
|
||
|
||
if not endpoint_url:
|
||
raise ValueError("缺少接口地址,请先选择提供方或手动填写 base_url")
|
||
if not llm_model_name:
|
||
raise ValueError("缺少模型名称,请填写 llm_model_name")
|
||
if provider not in {"ollama"} and not api_key:
|
||
raise ValueError("缺少 API Key,请填写后再测试")
|
||
|
||
timeout = int(payload.get("llm_timeout") or 120)
|
||
temperature = float(payload.get("llm_temperature") or 0.7)
|
||
top_p = float(payload.get("llm_top_p") or 0.9)
|
||
max_tokens = int(payload.get("llm_max_tokens") or 2048)
|
||
system_prompt = (payload.get("llm_system_prompt") or "").strip()
|
||
|
||
started_at = time.perf_counter()
|
||
preview = await asyncio.to_thread(
|
||
cls._send_test_request,
|
||
provider_meta.get("protocol", "openai_compatible"),
|
||
provider,
|
||
endpoint_url,
|
||
api_key,
|
||
llm_model_name,
|
||
timeout,
|
||
temperature,
|
||
top_p,
|
||
max_tokens,
|
||
system_prompt,
|
||
)
|
||
latency_ms = int((time.perf_counter() - started_at) * 1000)
|
||
|
||
return {
|
||
"provider": provider,
|
||
"endpoint_url": endpoint_url,
|
||
"llm_model_name": llm_model_name,
|
||
"latency_ms": latency_ms,
|
||
"preview": preview[:200],
|
||
}
|
||
|
||
@classmethod
|
||
def _send_test_request(
|
||
cls,
|
||
protocol: str,
|
||
provider: Optional[str],
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
llm_model_name: str,
|
||
timeout: int,
|
||
temperature: float,
|
||
top_p: float,
|
||
max_tokens: int,
|
||
system_prompt: str,
|
||
) -> str:
|
||
if protocol == "anthropic":
|
||
return cls._test_anthropic(
|
||
endpoint_url,
|
||
api_key,
|
||
llm_model_name,
|
||
timeout,
|
||
temperature,
|
||
top_p,
|
||
max_tokens,
|
||
system_prompt,
|
||
)
|
||
|
||
if protocol == "gemini":
|
||
return cls._test_gemini(
|
||
endpoint_url,
|
||
api_key,
|
||
llm_model_name,
|
||
timeout,
|
||
temperature,
|
||
top_p,
|
||
max_tokens,
|
||
system_prompt,
|
||
)
|
||
|
||
return cls._test_openai_compatible(
|
||
provider,
|
||
endpoint_url,
|
||
api_key,
|
||
llm_model_name,
|
||
timeout,
|
||
temperature,
|
||
top_p,
|
||
max_tokens,
|
||
system_prompt,
|
||
)
|
||
|
||
@classmethod
|
||
def _test_openai_compatible(
|
||
cls,
|
||
provider: Optional[str],
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
llm_model_name: str,
|
||
timeout: int,
|
||
temperature: float,
|
||
top_p: float,
|
||
max_tokens: int,
|
||
system_prompt: str,
|
||
) -> str:
|
||
url = cls._join_endpoint(endpoint_url, "/chat/completions")
|
||
headers = {
|
||
"Content-Type": "application/json",
|
||
}
|
||
if api_key:
|
||
headers["Authorization"] = f"Bearer {api_key}"
|
||
|
||
payload = {
|
||
"model": llm_model_name,
|
||
"messages": cls._build_openai_messages(system_prompt),
|
||
"temperature": temperature,
|
||
"top_p": top_p,
|
||
"max_tokens": min(max_tokens, 256),
|
||
"stream": False,
|
||
}
|
||
response = cls._request_json(url, headers, payload, timeout)
|
||
choices = response.get("choices") or []
|
||
if not choices:
|
||
raise ValueError("测试请求已发送,但未收到模型返回内容")
|
||
|
||
content = choices[0].get("message", {}).get("content", "")
|
||
if isinstance(content, list):
|
||
content = "".join(
|
||
item.get("text", "") if isinstance(item, dict) else str(item)
|
||
for item in content
|
||
)
|
||
content = str(content).strip()
|
||
if not content:
|
||
raise ValueError("模型返回成功,但内容为空")
|
||
return content
|
||
|
||
@classmethod
|
||
def _test_anthropic(
|
||
cls,
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
llm_model_name: str,
|
||
timeout: int,
|
||
temperature: float,
|
||
top_p: float,
|
||
max_tokens: int,
|
||
system_prompt: str,
|
||
) -> str:
|
||
url = cls._join_endpoint(endpoint_url, "/v1/messages")
|
||
headers = {
|
||
"Content-Type": "application/json",
|
||
"x-api-key": api_key,
|
||
"anthropic-version": "2023-06-01",
|
||
}
|
||
payload = {
|
||
"model": llm_model_name,
|
||
"max_tokens": min(max_tokens, 256),
|
||
"temperature": temperature,
|
||
"top_p": top_p,
|
||
"messages": [
|
||
{
|
||
"role": "user",
|
||
"content": "请只回复“连接测试成功”。",
|
||
}
|
||
],
|
||
}
|
||
if system_prompt:
|
||
payload["system"] = system_prompt
|
||
|
||
response = cls._request_json(url, headers, payload, timeout)
|
||
content = response.get("content") or []
|
||
texts = []
|
||
for item in content:
|
||
if isinstance(item, dict) and item.get("type") == "text":
|
||
texts.append(item.get("text", ""))
|
||
preview = "".join(texts).strip()
|
||
if not preview:
|
||
raise ValueError("模型返回成功,但内容为空")
|
||
return preview
|
||
|
||
@classmethod
|
||
def _test_gemini(
|
||
cls,
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
llm_model_name: str,
|
||
timeout: int,
|
||
temperature: float,
|
||
top_p: float,
|
||
max_tokens: int,
|
||
system_prompt: str,
|
||
) -> str:
|
||
model_path = llm_model_name if llm_model_name.startswith("models/") else f"models/{llm_model_name}"
|
||
encoded_model_path = "/".join(urllib.parse.quote(part) for part in model_path.split("/"))
|
||
url = f"{endpoint_url.rstrip('/')}/{encoded_model_path}:generateContent?key={urllib.parse.quote(api_key)}"
|
||
headers = {
|
||
"Content-Type": "application/json",
|
||
}
|
||
payload = {
|
||
"contents": [
|
||
{
|
||
"role": "user",
|
||
"parts": [{"text": "请只回复“连接测试成功”。"}],
|
||
}
|
||
],
|
||
"generationConfig": {
|
||
"temperature": temperature,
|
||
"topP": top_p,
|
||
"maxOutputTokens": min(max_tokens, 256),
|
||
},
|
||
}
|
||
if system_prompt:
|
||
payload["systemInstruction"] = {
|
||
"parts": [{"text": system_prompt}],
|
||
}
|
||
|
||
response = cls._request_json(url, headers, payload, timeout)
|
||
candidates = response.get("candidates") or []
|
||
if not candidates:
|
||
raise ValueError("测试请求已发送,但未收到模型返回内容")
|
||
|
||
parts = candidates[0].get("content", {}).get("parts", [])
|
||
preview = "".join(
|
||
part.get("text", "") for part in parts if isinstance(part, dict)
|
||
).strip()
|
||
if not preview:
|
||
raise ValueError("模型返回成功,但内容为空")
|
||
return preview
|
||
|
||
@staticmethod
|
||
def _build_openai_messages(system_prompt: str) -> List[Dict[str, str]]:
|
||
messages: List[Dict[str, str]] = []
|
||
if system_prompt:
|
||
messages.append({"role": "system", "content": system_prompt})
|
||
messages.append({"role": "user", "content": "请只回复“连接测试成功”。"})
|
||
return messages
|
||
|
||
@staticmethod
|
||
def _join_endpoint(base_url: str, suffix: str) -> str:
|
||
normalized = base_url.rstrip("/")
|
||
if normalized.endswith(suffix):
|
||
return normalized
|
||
return f"{normalized}{suffix}"
|
||
|
||
@classmethod
|
||
def _request_json(
|
||
cls,
|
||
url: str,
|
||
headers: Dict[str, str],
|
||
payload: Dict[str, Any],
|
||
timeout: int,
|
||
) -> Dict[str, Any]:
|
||
request = urllib.request.Request(
|
||
url,
|
||
data=json.dumps(payload).encode("utf-8"),
|
||
headers=headers,
|
||
method="POST",
|
||
)
|
||
|
||
ctx = None
|
||
if os.getenv("DISABLE_SSL_VERIFY", "").lower() in ("1", "true", "yes"):
|
||
ctx = ssl.create_default_context()
|
||
ctx.check_hostname = False
|
||
ctx.verify_mode = ssl.CERT_NONE
|
||
|
||
try:
|
||
with urllib.request.urlopen(request, timeout=timeout, context=ctx) as response:
|
||
body = response.read().decode("utf-8")
|
||
if not body:
|
||
return {}
|
||
return json.loads(body)
|
||
except urllib.error.HTTPError as exc:
|
||
error_body = exc.read().decode("utf-8", errors="ignore")
|
||
message = cls._extract_error_message(error_body) or error_body[:300] or str(exc)
|
||
raise ValueError(f"模型测试失败(HTTP {exc.code}):{message}") from exc
|
||
except urllib.error.URLError as exc:
|
||
reason = exc.reason
|
||
if isinstance(reason, socket.timeout):
|
||
raise ValueError("模型测试超时,请检查网络或调大超时时间") from exc
|
||
raise ValueError(f"模型测试失败:{reason}") from exc
|
||
except socket.timeout as exc:
|
||
raise ValueError("模型测试超时,请检查网络或调大超时时间") from exc
|
||
except json.JSONDecodeError as exc:
|
||
raise ValueError("模型服务返回了无法解析的响应,请检查接口地址是否正确") from exc
|
||
|
||
@staticmethod
|
||
def _extract_error_message(error_body: str) -> Optional[str]:
|
||
if not error_body:
|
||
return None
|
||
|
||
try:
|
||
payload = json.loads(error_body)
|
||
except json.JSONDecodeError:
|
||
return error_body.strip()
|
||
|
||
if isinstance(payload, dict):
|
||
if isinstance(payload.get("error"), dict):
|
||
return payload["error"].get("message") or payload["error"].get("type")
|
||
return payload.get("message") or payload.get("detail")
|
||
|
||
return None
|
||
|
||
# ==================== 对话文本生成(RAG 调用) ====================
|
||
|
||
@classmethod
|
||
async def generate_text(
|
||
cls,
|
||
provider: Optional[str],
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
llm_model_name: str,
|
||
timeout: int,
|
||
temperature: float,
|
||
top_p: float,
|
||
max_tokens: int,
|
||
system_prompt: str,
|
||
messages: List[Dict[str, Any]],
|
||
) -> str:
|
||
"""根据多轮对话消息生成回复文本(供知识库 RAG 使用)"""
|
||
provider_meta = PROVIDER_MAP.get(provider, PROVIDER_MAP["custom"])
|
||
endpoint_url = (endpoint_url or provider_meta.get("default_endpoint_url") or "").strip()
|
||
llm_model_name = (llm_model_name or "").strip()
|
||
api_key = (api_key or "").strip()
|
||
|
||
if not endpoint_url:
|
||
raise ValueError("缺少接口地址,请先选择提供方或手动填写 base_url")
|
||
if not llm_model_name:
|
||
raise ValueError("缺少模型名称,请填写 llm_model_name")
|
||
if provider not in {"ollama"} and not api_key:
|
||
raise ValueError("缺少 API Key,请填写后再调用")
|
||
|
||
normalized = cls._normalize_messages(messages)
|
||
if not normalized:
|
||
raise ValueError("缺少有效的对话消息")
|
||
|
||
protocol = provider_meta.get("protocol", "openai_compatible")
|
||
return await asyncio.to_thread(
|
||
cls._chat_completion,
|
||
protocol,
|
||
endpoint_url,
|
||
api_key,
|
||
llm_model_name,
|
||
int(timeout or 120),
|
||
float(temperature or 0.7),
|
||
float(top_p or 0.9),
|
||
int(max_tokens or 2048),
|
||
(system_prompt or "").strip(),
|
||
normalized,
|
||
)
|
||
|
||
@staticmethod
|
||
def _normalize_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, str]]:
|
||
normalized: List[Dict[str, str]] = []
|
||
for item in messages or []:
|
||
if not isinstance(item, dict):
|
||
continue
|
||
role = str(item.get("role") or "").strip().lower()
|
||
if role not in {"system", "user", "assistant"}:
|
||
continue
|
||
content = item.get("content", "")
|
||
if isinstance(content, list):
|
||
content = "".join(
|
||
part.get("text", "") if isinstance(part, dict) else str(part)
|
||
for part in content
|
||
)
|
||
content = str(content).strip()
|
||
if not content:
|
||
continue
|
||
normalized.append({"role": role, "content": content})
|
||
return normalized
|
||
|
||
@classmethod
|
||
def _chat_completion(
|
||
cls,
|
||
protocol: str,
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
llm_model_name: str,
|
||
timeout: int,
|
||
temperature: float,
|
||
top_p: float,
|
||
max_tokens: int,
|
||
system_prompt: str,
|
||
messages: List[Dict[str, str]],
|
||
) -> str:
|
||
if protocol == "anthropic":
|
||
return cls._chat_anthropic(
|
||
endpoint_url, api_key, llm_model_name, timeout,
|
||
temperature, top_p, max_tokens, system_prompt, messages,
|
||
)
|
||
if protocol == "gemini":
|
||
return cls._chat_gemini(
|
||
endpoint_url, api_key, llm_model_name, timeout,
|
||
temperature, top_p, max_tokens, system_prompt, messages,
|
||
)
|
||
return cls._chat_openai_compatible(
|
||
endpoint_url, api_key, llm_model_name, timeout,
|
||
temperature, top_p, max_tokens, system_prompt, messages,
|
||
)
|
||
|
||
@classmethod
|
||
def _chat_openai_compatible(
|
||
cls,
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
llm_model_name: str,
|
||
timeout: int,
|
||
temperature: float,
|
||
top_p: float,
|
||
max_tokens: int,
|
||
system_prompt: str,
|
||
messages: List[Dict[str, str]],
|
||
) -> str:
|
||
url = cls._join_endpoint(endpoint_url, "/chat/completions")
|
||
headers = {"Content-Type": "application/json"}
|
||
if api_key:
|
||
headers["Authorization"] = f"Bearer {api_key}"
|
||
|
||
payload_messages: List[Dict[str, str]] = []
|
||
if system_prompt:
|
||
payload_messages.append({"role": "system", "content": system_prompt})
|
||
payload_messages.extend(messages)
|
||
|
||
payload = {
|
||
"model": llm_model_name,
|
||
"messages": payload_messages,
|
||
"temperature": temperature,
|
||
"top_p": top_p,
|
||
"max_tokens": max_tokens,
|
||
"stream": False,
|
||
}
|
||
response = cls._request_json(url, headers, payload, timeout)
|
||
choices = response.get("choices") or []
|
||
if not choices:
|
||
raise ValueError("请求已发送,但未收到模型返回内容")
|
||
content = choices[0].get("message", {}).get("content", "")
|
||
if isinstance(content, list):
|
||
content = "".join(
|
||
item.get("text", "") if isinstance(item, dict) else str(item)
|
||
for item in content
|
||
)
|
||
content = str(content).strip()
|
||
if not content:
|
||
raise ValueError("模型返回成功,但内容为空")
|
||
return content
|
||
|
||
@classmethod
|
||
def _chat_anthropic(
|
||
cls,
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
llm_model_name: str,
|
||
timeout: int,
|
||
temperature: float,
|
||
top_p: float,
|
||
max_tokens: int,
|
||
system_prompt: str,
|
||
messages: List[Dict[str, str]],
|
||
) -> str:
|
||
url = cls._join_endpoint(endpoint_url, "/v1/messages")
|
||
headers = {
|
||
"Content-Type": "application/json",
|
||
"x-api-key": api_key,
|
||
"anthropic-version": "2023-06-01",
|
||
}
|
||
payload = {
|
||
"model": llm_model_name,
|
||
"max_tokens": max_tokens,
|
||
"temperature": temperature,
|
||
"top_p": top_p,
|
||
"messages": [
|
||
{"role": m["role"], "content": [{"type": "text", "text": m["content"]}]}
|
||
for m in messages
|
||
if m["role"] in {"user", "assistant"}
|
||
],
|
||
}
|
||
if system_prompt:
|
||
payload["system"] = system_prompt
|
||
|
||
response = cls._request_json(url, headers, payload, timeout)
|
||
content = response.get("content") or []
|
||
texts = [
|
||
item.get("text", "")
|
||
for item in content
|
||
if isinstance(item, dict) and item.get("type") == "text"
|
||
]
|
||
preview = "".join(texts).strip()
|
||
if not preview:
|
||
raise ValueError("模型返回成功,但内容为空")
|
||
return preview
|
||
|
||
@classmethod
|
||
def _chat_gemini(
|
||
cls,
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
llm_model_name: str,
|
||
timeout: int,
|
||
temperature: float,
|
||
top_p: float,
|
||
max_tokens: int,
|
||
system_prompt: str,
|
||
messages: List[Dict[str, str]],
|
||
) -> str:
|
||
model_path = llm_model_name if llm_model_name.startswith("models/") else f"models/{llm_model_name}"
|
||
encoded_model_path = "/".join(urllib.parse.quote(part) for part in model_path.split("/"))
|
||
url = f"{endpoint_url.rstrip('/')}/{encoded_model_path}:generateContent?key={urllib.parse.quote(api_key)}"
|
||
headers = {"Content-Type": "application/json"}
|
||
|
||
role_map = {"user": "user", "assistant": "model"}
|
||
contents = [
|
||
{"role": role_map[m["role"]], "parts": [{"text": m["content"]}]}
|
||
for m in messages
|
||
if m["role"] in role_map
|
||
]
|
||
payload = {
|
||
"contents": contents,
|
||
"generationConfig": {
|
||
"temperature": temperature,
|
||
"topP": top_p,
|
||
"maxOutputTokens": max_tokens,
|
||
},
|
||
}
|
||
if system_prompt:
|
||
payload["systemInstruction"] = {"parts": [{"text": system_prompt}]}
|
||
|
||
response = cls._request_json(url, headers, payload, timeout)
|
||
candidates = response.get("candidates") or []
|
||
if not candidates:
|
||
raise ValueError("请求已发送,但未收到模型返回内容")
|
||
parts = candidates[0].get("content", {}).get("parts", [])
|
||
preview = "".join(
|
||
part.get("text", "") for part in parts if isinstance(part, dict)
|
||
).strip()
|
||
if not preview:
|
||
raise ValueError("模型返回成功,但内容为空")
|
||
return preview
|
||
|
||
# ==================== Embedding 向量生成与测试 ====================
|
||
|
||
@classmethod
|
||
async def generate_embedding(
|
||
cls,
|
||
provider: Optional[str],
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
llm_model_name: str,
|
||
text: str,
|
||
timeout: int = 60,
|
||
dimension: Optional[int] = None,
|
||
) -> List[float]:
|
||
"""调用 OpenAI 兼容的 /embeddings 接口生成单条文本向量"""
|
||
provider_meta = PROVIDER_MAP.get(provider, PROVIDER_MAP["custom"])
|
||
endpoint_url = (endpoint_url or provider_meta.get("default_endpoint_url") or "").strip()
|
||
llm_model_name = (llm_model_name or "").strip()
|
||
api_key = (api_key or "").strip()
|
||
text = (text or "").strip()
|
||
|
||
if not endpoint_url:
|
||
raise ValueError("缺少接口地址,请先选择提供方或手动填写 base_url")
|
||
if not llm_model_name:
|
||
raise ValueError("缺少模型名称,请填写 embedding 模型标识")
|
||
if provider not in {"ollama"} and not api_key:
|
||
raise ValueError("缺少 API Key,请填写后再调用")
|
||
if not text:
|
||
raise ValueError("待向量化文本为空")
|
||
|
||
vectors = await asyncio.to_thread(
|
||
cls._embeddings_request,
|
||
endpoint_url,
|
||
api_key,
|
||
llm_model_name,
|
||
[text[:8000]],
|
||
int(timeout or 60),
|
||
dimension,
|
||
)
|
||
if not vectors:
|
||
raise ValueError("Embedding 接口返回为空")
|
||
return vectors[0]
|
||
|
||
@classmethod
|
||
def _embeddings_request(
|
||
cls,
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
llm_model_name: str,
|
||
inputs: List[str],
|
||
timeout: int,
|
||
dimension: Optional[int],
|
||
) -> List[List[float]]:
|
||
url = cls._join_endpoint(endpoint_url, "/embeddings")
|
||
headers = {"Content-Type": "application/json"}
|
||
if api_key:
|
||
headers["Authorization"] = f"Bearer {api_key}"
|
||
|
||
payload: Dict[str, Any] = {
|
||
"model": llm_model_name,
|
||
"input": inputs,
|
||
}
|
||
if dimension:
|
||
payload["dimensions"] = dimension
|
||
|
||
response = cls._request_json(url, headers, payload, timeout)
|
||
data = response.get("data") or []
|
||
if not data:
|
||
raise ValueError("Embedding 请求已发送,但未收到向量数据")
|
||
|
||
vectors: List[List[float]] = []
|
||
for item in data:
|
||
embedding = item.get("embedding") if isinstance(item, dict) else None
|
||
if embedding:
|
||
vectors.append([float(x) for x in embedding])
|
||
if not vectors:
|
||
raise ValueError("Embedding 返回结果中未解析到向量")
|
||
return vectors
|
||
|
||
@classmethod
|
||
async def test_embedding_connection(cls, payload: Dict[str, Any]) -> Dict[str, Any]:
|
||
"""测试 embedding 模型连通性,并返回实际向量维度"""
|
||
provider = payload.get("provider")
|
||
provider_meta = PROVIDER_MAP.get(provider, PROVIDER_MAP["custom"])
|
||
endpoint_url = (payload.get("endpoint_url") or provider_meta.get("default_endpoint_url") or "").strip()
|
||
llm_model_name = (payload.get("llm_model_name") or "").strip()
|
||
api_key = (payload.get("api_key") or "").strip()
|
||
timeout = int(payload.get("llm_timeout") or 60)
|
||
dimension = payload.get("embedding_dimension")
|
||
|
||
started_at = time.perf_counter()
|
||
vector = await cls.generate_embedding(
|
||
provider,
|
||
endpoint_url,
|
||
api_key,
|
||
llm_model_name,
|
||
"连接测试",
|
||
timeout=timeout,
|
||
dimension=int(dimension) if dimension else None,
|
||
)
|
||
latency_ms = int((time.perf_counter() - started_at) * 1000)
|
||
|
||
return {
|
||
"provider": provider,
|
||
"endpoint_url": endpoint_url,
|
||
"llm_model_name": llm_model_name,
|
||
"latency_ms": latency_ms,
|
||
"dimension": len(vector),
|
||
"preview": f"成功生成 {len(vector)} 维向量",
|
||
}
|
||
|