54 lines
2.3 KiB
Python
54 lines
2.3 KiB
Python
"""Model choices and safe, task-specific extraction settings."""
|
||
from __future__ import annotations
|
||
|
||
import re
|
||
|
||
|
||
def catalog_for(database):
|
||
from model_gateway import Catalog
|
||
catalog = Catalog(database.path)
|
||
catalog.refresh()
|
||
return catalog
|
||
|
||
|
||
def model_choices(database):
|
||
catalog = catalog_for(database)
|
||
answers, _, _, _ = catalog.plan()
|
||
default = answers[0].id if answers else ''
|
||
return [{"id": outlet.id, "name": outlet.name, "model": outlet.config.get("model", ""),
|
||
"kind": outlet.kind, "default": outlet.id == default}
|
||
for outlet in catalog.outlets.values() if outlet.kind != 'comfyui']
|
||
|
||
|
||
def resolve_model(database, options):
|
||
catalog = catalog_for(database)
|
||
provider = options.get('model_provider_id', '')
|
||
if provider:
|
||
outlet = catalog.outlets.get(provider)
|
||
else:
|
||
answers, _, _, _ = catalog.plan()
|
||
outlet = answers[0] if answers else None
|
||
if outlet is None or outlet.kind == 'comfyui':
|
||
raise ValueError('所选模型不可用或已停用,请选择已启用的文本模型')
|
||
options.update(model_provider_id=outlet.id, model_name=outlet.name,
|
||
model_timeout_seconds=options.get('model_timeout_seconds', 90))
|
||
return outlet
|
||
|
||
|
||
def safe_model_error(error):
|
||
"""Never persist provider response bodies, URLs or credentials."""
|
||
value = str(error).lower()
|
||
if any(word in value for word in ('timeout', '时限', '超时')):
|
||
return '模型请求超时,请增加单次等待时间或更换模型后继续'
|
||
code = re.search(r'上游返回\s+(\d{3})', value)
|
||
if code:
|
||
status = code.group(1)
|
||
advice = {'401': '请检查模型密钥', '403': '请检查模型授权',
|
||
'429': '模型限流或额度不足,请稍后重试或更换模型'}.get(status, '请检查模型服务后重试或更换模型')
|
||
return f'模型服务返回 HTTP {status};{advice}'
|
||
if any(word in value for word in ('connect', 'network', 'transport')):
|
||
return '无法连接模型服务,请检查网络和模型地址后重试'
|
||
if '熔断' in value:
|
||
return '模型连续失败,暂时不可用,请稍后重试或更换模型'
|
||
return '模型服务未返回有效结果,请检查模型配置或更换模型'
|